KingOfEgg
首页项目归档照片墙音乐灵境说说杂谈友链关于
封面

Softmax优化(一) 从Naive到两遍融合

写作时间:2026-08-17 12:00:00
# CUDA
# Softmax
# GPU优化
# 性能调优

Softmax优化:从Naive到两遍融合

测试平台:RTX 5080 (Blackwell, sm_120),M=N=8192,理论带宽960 GB/s

前言:softmax是attention的核心算子,也是典型memory-bound算子。深度学习框架中softmax算子几乎是必不可少的组件,在实现时一般也是GOU kernel function形式。

| 版本 | 核心优化               | 配置         | 带宽(GB/s) | 相对v1 |
| --- | ---------------------- | ------------ | ---------- | ------- |
| v1  | naive:一个线程一行      | 1×256        | ~150       | 基准    |
| v2  | shared memory树形归约  | 256×1        | ~820       | ~5.5×   |
| v3  | warp shuffle两级归约   | 256×1        | ~810       | 持平    |
| v4  | float4向量化加载       | 256×1        | ~820       | 略升    |
| v5  | online两遍融合         | 256×1        | ~800       | 持平    |

v2起各版本都在~85%峰值附近,真正的瓶颈可能是DRAM带宽,无法再突破DRAM上限。但5080显卡无dram读取计数器,无法获取真实带宽利用率。

1 Safe Softmax与三遍扫描

y_i = exp(x_i - max) / Σ exp(x_j - max)

三遍扫描:

  1. 第一遍求行最大值 max
  2. 第二遍求指数和 sum
  3. 第三遍归一化写出 y

2 v1版本:朴素实现

一个线程负责一行,串行扫三遍:

__global__ void softmax_v1_naive(float *input, float *output, int M, int N) {
    int row = blockIdx.y * blockDim.y + threadIdx.y;
    if (row >= M) return;
    float *x = input + row * N;
    float *y = output + row * N;
    // 求最大值
    float max = -INFINITY;
    for (int k = 0; k < N; k++) max = fmax(max, x[k]);
    // 指数求和
    float sum = 0.0f;
    for (int k = 0; k < N; k++) sum += expf(x[k] - max);
    // 归一化
    float inv_sum = 1.0f / sum;
    for (int k = 0; k < N; k++) y[k] = expf(x[k] - max) * inv_sum;
}

结果~150 GB/s,只有峰值的15%。

分析:

  1. 并行度太低:一个线程要串行扫完整行8192个元素,8192次循环。
  2. 非合并访存:相邻线程(threadIdx.y相邻)处理的是不同行,它们访问的地址相差N(32KB),一个warp的32次load落在32个不同cache line上,非合并访存。

3 v2版本:一个block一行 + shared memory归约

让一个block(256线程)协作处理一行,每个线程只负责一行的一部分,再在block内做归约。归约用shared memory做树形规约:

__global__ void softmax_v2_sharedreduce(float *input, float *output, int M, int N) {
    int row = blockIdx.y;   // 一个block处理一行
    int tx = threadIdx.x;
    if (row >= M) return;
    float *x = input + row * N;
    float *y = output + row * N;
    extern __shared__ float smem[];

    // Pass 1:并行求最大值
    float max_val = -INFINITY;
    for (int i = tx; i < N; i += blockDim.x) max_val = fmaxf(max_val, x[i]);
    smem[tx] = max_val;
    __syncthreads();
    for (int s = blockDim.x / 2; s > 0; s >>= 1) {
        if (tx < s) smem[tx] = fmaxf(smem[tx], smem[tx + s]);
        __syncthreads();
    }
    max_val = smem[0];
    __syncthreads();

    // Pass 2 / Pass 3:求sum、归一化写出,结构同上
}

结果跳到~820 GB/s,是v1的5.5倍。

分析:

  1. 每个线程用 for (i = tx; i < N; i += blockDim.x) 跨步访问,256个线程把一行铺满,访存变成合并的。
  2. 归约用树形 for (s = blockDim.x/2; s > 0; s >>= 1),256 → 1 只需要8步,每步一次 __syncthreads()。

如果sum归约的循环条件写成 for (int i = blockDim.x/2; i < N; i >>= 1),一方面 i < N永远成立会死循环,另一方面 smem[tx + i]会越界。正确写法是上面的 s >>= 1且 s > 0且 if (tx < s)保护。

4 v3版本:warp shuffle减少同步

v2的树形归约每一步都要 __syncthreads()(整个block同步8次),而warp内的32个线程是天然同步的。warp内用 __shfl_down_sync 归约(寄存器级、免同步),warp之间再用一小块shared memory归约,同步次数从8次降到2次。

// warp内归约:5次shuffle,无同步
__device__ float warpReduceMax(float val) {
    for (int offset = 16; offset > 0; offset >>= 1)
        val = fmaxf(val, __shfl_down_sync(0xffffffff, val, offset));
    return val;
}

// block归约:warp内shuffle + warp间smem
__device__ float blockReduceMax(float val) {
    __shared__ float warp_max[32];
    int lane = threadIdx.x & 31;
    int wid  = threadIdx.x >> 5;
    val = warpReduceMax(val);
    if (lane == 0) warp_max[wid] = val;
    __syncthreads();                              // 仅1次同步收集各warp
    int nwarps = blockDim.x / 32;
    val = (lane < nwarps) ? warp_max[lane] : -INFINITY;
    if (wid == 0) val = warpReduceMax(val);       // 第0个warp收尾
    __shared__ float result;
    if (threadIdx.x == 0) result = val;
    __syncthreads();
    return result;
}

结果~810 GB/s,和v2持平。

5 v4版本:float4向量化加载

既然瓶颈在访存,就把load/store指令数降下来:一次读4个float(128-bit的 float4),三遍循环全部向量化。

const float4* x4 = reinterpret_cast<const float4*>(input + row * N);
float4*       y4 = reinterpret_cast<float4*>(output + row * N);
int N4 = N / 4;   // 计数单位从N变成N/4

// Pass 1:一次读4个,取4分量最大值
float max_val = -INFINITY;
for (int i = tx; i < N4; i += blockDim.x) {
    float4 v = x4[i];
    max_val = fmaxf(max_val, fmaxf(fmaxf(v.x, v.y), fmaxf(v.z, v.w)));
}
// Pass 2 / Pass 3同理,4个分量分别exp后累加 / 构造float4一次写出

分析:

  1. load/store指令从4条 LDG.32 变成1条 LDG.128,指令数大幅下降。
  2. 4个连续线程访问4个连续float4,正好凑满128字节cache line,合并度更好。

结果~820 GB/s,相对v3略升。收益有限,因为DRAM已经接近满负荷,指令数减少只能缓解,无法突破带宽上限。

6 v5版本:两遍融合(online softmax)

三遍扫描读input三次,能不能把max和sum合并到一次读取里,只读两次?这就是online softmax的思路。

问题在于:sum依赖最终的max,而max是扫到最后才知道的。解决办法是维护一个running (m, s),遇到更大的值就把旧的sum按指数缩放:

遇到新元素x,设 m_new = max(m, x):
    s = s * exp(m - m_new) + exp(x - m_new)
    m = m_new

当 x <= m 时 exp(m - m_new) = 1,等价于 s += exp(x - m);当 x > m 时,旧sum乘以 exp(m - x) 缩放到新基准。

block内每个线程先算自己的局部(m, s),归约时先归约出全局max,再把每个线程的s缩放到全局max后求和:

// Pass 1:一遍读取,在线更新(m, s)
float m = -INFINITY, s = 0.0f;
for (int i = tx; i < N4; i += blockDim.x) {
    float4 v = x4[i];
    float mn;
    mn = fmaxf(m, v.x); s = s * expf(m - mn) + expf(v.x - mn); m = mn;
    mn = fmaxf(m, v.y); s = s * expf(m - mn) + expf(v.y - mn); m = mn;
    mn = fmaxf(m, v.z); s = s * expf(m - mn) + expf(v.z - mn); m = mn;
    mn = fmaxf(m, v.w); s = s * expf(m - mn) + expf(v.w - mn); m = mn;
}
float global_m = blockReduceMax(m);
s = s * expf(m - global_m);   // 缩放到全局max
float global_s = blockReduceSum(s);

// Pass 2:一遍读取,归一化写出

结果~800 GB/s,没有比v4快。一开始我以为两遍融合能省掉一次DRAM读,应该快1/3,结果没变,于是用ncu查了真实数据。

7 用ncu验证:瓶颈是DRAM带宽

ncu数据(8192×8192,1 sector = 32字节,读一遍 = 256 MiB = 8,388,608 sectors):

v4(三遍扫描):

gpu__time_duration.sum               us       622.88
lts__t_bytes.sum                  Gbyte         1.08
lts__t_sectors_op_read.sum       sector   25,331,031   ≈ 3次
lts__t_sectors_op_write.sum      sector    8,388,623   ≈ 1次
dram__bytes_read.sum                         (!) n/a
dram__bytes_write.sum                        (!) n/a

v5(两遍融合):

gpu__time_duration.sum               us       631.20
lts__t_bytes.sum                  Mbyte       820.83
lts__t_sectors_op_read.sum       sector   17,260,927   ≈ 2次
lts__t_sectors_op_write.sum      sector    8,388,620   ≈ 1次
dram__bytes_read.sum                         (!) n/a
dram__bytes_write.sum                        (!) n/a
  • gpu__time_duration.sum:kernel执行时长(微秒)。
  • lts__t_sectors_op_read/write.sum:到达L2缓存的读/写请求量,单位sector(32字节)。lts即L2 slice。
  • dram__bytes_read/write.sum:DRAM流量,GeForce消费卡不开放这个计数器,所以是n/a。

对比:

  1. L2读流量:v4是v5的1.47倍(25.3M vs 17.3M sectors),证明v5确实省掉了一次读。
  2. 但两者时间几乎一样(622 vs 631 us)。

结论:省掉的那次读命中了L2,根本没打到DRAM。一行32KB,同一个block里三个pass间隔极短,数据一直躺在L2里。所以三遍扫描的DRAM真实流量早就已经是下限「读1次 + 写1次 = 512 MiB」,两遍融合省的是L2带宽(L2不紧张),不省DRAM,因此不加速。

用时间反推真实DRAM带宽:

512 MiB / 622.88 us ≈ 862 GB/s   (约89%峰值)

8 总结与相关参考

参考博客:

深入浅出 GPU优化系列:softmax

AIInfraGuide:Softmax优化

特别鸣谢:Deepseek-v4-Pro、Trae Agent

avatar

KingOfEgg

otaku change the world

RECOMMENDED

GEMM 优化(一) Cuda Core计算

2026-08-11 12:00:00

CUBLAS GEMM 函数用法

2026-08-12 12:00:00

GEMM优化(三)大矩阵适配和极致优化

2026-09-03 12:00:00

Table of Contents